8. Neural Network Ex.2 Multi-Class Classification
When a problem requires us to distinguish between more than two categories, we use multi-class classification. Real-world examples include:
- Image Recognition: Classifying an image as a Cat (0), Dog (1), or Bird (2).
- Digit Recognition: The MNIST dataset, predicting digits 0 through 9.
If we have 3 possible classes, our target labels will be integers representing the classes:
Class 0
Class 1
Class 2
The Output Layer
In multi-class classification, the neural network's final layer must output one logit per class. For a 3-class problem, the model produces a vector of 3 logits:
[z1, z2, z3]
The Softmax Function
Just like we used Sigmoid for binary classification, we use the Softmax function for multi-class problems. Softmax converts the vector of raw logits into a Probability Distribution, meaning every value is between 0 and 1, and the total sum of all values is exactly 1.0.
Try playing with the Softmax Visualizer below! You can change the raw logits and see how Softmax normalizes them into probabilities. The Temperature slider controls how "confident" the distribution is (a feature heavily used in Large Language Models).
Interactive Softmax & Temperature Visualizer
Adjust the raw Logits or move the Temperature slider to explore how temperature controls randomness in LLM next-token generation (low T makes output greedy/deterministic, high T flattens probabilities).
Cross-Entropy Loss
To train a multi-class model, we use Cross-Entropy Loss. It looks at the probability assigned to the correct class and penalizes the model if that probability is low.
In PyTorch, the standard loss function is CrossEntropyLoss. Crucially, it expects raw logits as input, not probabilities! Under the hood, PyTorch efficiently combines LogSoftmax and Negative Log-Likelihood Loss (NLLLoss).
loss_function = torch.nn.CrossEntropyLoss()
Making Predictions
To get the final class prediction from the network's output, we don't necessarily need to compute the exact probabilities. Because the exponential function is monotonic, the class with the highest raw logit will also have the highest probability.
We just find the index of the maximum logit using argmax:
prediction = logits.argmax(dim=1)
Complete Runnable Example
Let's build and train a multi-class network! We'll create a dummy dataset with 3 distinct clusters.